6. 理解 GQA
GQA(Grouped Query Attention,分组查询注意力)是 Google 在 2023 年论文《GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints》中提出的一种 Attention 优化方法。它的核心目标是:
在尽量不损失模型效果的前提下,大幅降低 KV Cache 的内存占用和推理成本。 (ACL Anthology)
如果你已经理解了 Transformer 的 Multi-Head Attention(MHA),那么 GQA 本质上就是:
让多个 Query Head 共享同一组 Key 和 Value Head。
先理解问题:为什么要优化 Attention?
标准 Transformer 的 Attention:
对于每个 Head,都有独立的:
- Query
- Key
- Value
假设:
hidden_size = 4096
num_heads = 32
那么:
Q: 32个头
K: 32个头
V: 32个头
结构如下:
Head1: Q1 K1 V1
Head2: Q2 K2 V2
...
Head32: Q32 K32 V32
这就是:
MHA (Multi Head Attention)
MHA 最大的问题
在训练阶段问题不大。
但在推理阶段(尤其生成长文本):
模型要保存历史 Token 的
K Cache
V Cache
例如:
上下文长度 = 128K
Head数 = 32
KV Cache 会非常巨大。
实际上:
大模型推理时最大的瓶颈之一就是 KV Cache。 (IBM)
第一个解决方案:MQA
Google 之前提出:
MQA
(Multi Query Attention)
思路极其简单:
保留多个 Query Head
Q1
Q2
...
Q32
但是:
所有 Head 共用一个 K 和 V
K_shared
V_shared
变成:
Q1 ─┐
Q2 ─┤
Q3 ─┤
... ├── K_shared
Q32─┘ V_shared
即:
32个Q
1个K
1个V
KV Cache 直接缩小:
理论上节省:
的 KV 存储。
推理速度暴涨。 (ACL Anthology)
但 MQA 有副作用
虽然快了:
Q1
Q2
Q3
...
Q32
都在使用同一个:
K
V
很多 Head 的表达能力被压缩了。
效果通常会下降:
MHA > GQA > MQA
论文发现:
MQA 推理速度很好,但模型质量会有明显损失。 (ACL Anthology)
GQA 的核心思想
Google 想:
MQA太极端
MHA太昂贵
于是取中间方案:
Grouped Query Attention
假设:
32个Query Head
不要:
32个KV(MHA)
也不要:
1个KV(MQA)
而是:
8个KV
例如:
32个Q
8个KV
每4个 Query Head 共用一个 KV Head。
结构变成:
Q1
Q2
Q3
Q4
↓
KV1
Q5
Q6
Q7
Q8
↓
KV2
...
Q29
Q30
Q31
Q32
↓
KV8
这就是:
Group Size = 4
数学表示:
设:
那么:
例如:
即:
4个Q共享1个KV
三种 Attention 对比
假设:
32个Query Head
MHA
32Q
32K
32V
结构:
Q1 -> K1,V1
Q2 -> K2,V2
...
Q32->K32,V32
KV Cache:
100%
效果:
最好
MQA
32Q
1K
1V
结构:
Q1
Q2
...
Q32
共享K,V
KV Cache:
1/32
效果:
下降明显
GQA
32Q
8K
8V
结构:
4个Q共享1个KV
KV Cache:
1/4
效果:
接近MHA
所以:
MHA ←→ GQA ←→ MQA
GQA 本质上是:
MHA 与 MQA 之间的折中方案
为什么 GQA 特别适合大模型?
因为推理时:
KV Cache 大小约为:
其中:
- L = 上下文长度
= KV Head 数量
所以:
如果:
32 KV Heads
→
8 KV Heads
则:
KV Cache:
对于:
32K
64K
128K
长上下文模型,
收益巨大。
这也是为什么:
- Llama 2 70B
- Llama 3
- Mistral
- Qwen 2
- DeepSeek
等现代大模型广泛采用 GQA。
从代码角度理解
传统 Attention:
Q.shape = [B, L, 32, d]
K.shape = [B, L, 32, d]
V.shape = [B, L, 32, d]
GQA:
Q.shape = [B, L, 32, d]
K.shape = [B, L, 8, d]
V.shape = [B, L, 8, d]
然后:
K = repeat_interleave(K, 4)
V = repeat_interleave(V, 4)
逻辑上扩展成:
32 heads
供 Attention 使用。
例如:
KV1 → Q1,Q2,Q3,Q4
KV2 → Q5,Q6,Q7,Q8
...
从信息论角度理解
你之前问过:
RoPE 那些频率设计是不是拍脑袋?
GQA 反而比 RoPE 更容易理解。
GQA 背后的观察其实是:
Query Head 差异很大
不同 Head 学习不同模式:
语法
实体
位置
代码
推理
因此:
Q需要保留多样性
K/V Head 差异没那么大
论文实验发现:
很多 Head 的 Key/Value 存在冗余。
因此:
Q保持32个
KV压缩到8个
性能损失很小。 (Hugging Face)
一句话总结
GQA(Grouped Query Attention)可以理解成:
保留大量 Query Head 的表达能力,让多个 Query Head 共享较少数量的 Key/Value Head,从而大幅减少 KV Cache,占用更少显存,同时保持接近 Multi-Head Attention 的效果。 (ACL Anthology)
如果后面你准备从零实现 Llama/Qwen,我还可以继续讲:
- GQA 的完整数学推导
- Llama3 中
num_heads=32, num_kv_heads=8的具体实现 - PyTorch 版 GQA 源码逐行解析
- KV Cache 为什么能从 GQA 中获得巨大收益
- FlashAttention 与 GQA 的关系。